import triton
import triton.language as tl
import torch

@triton.jit
def slow_decode_kernel(
    q_ptr,
    k_cache_ptr,
    v_cache_ptr,
    output_ptr,
    cache_seqlens_ptr,
    block_table_ptr,
    num_heads: tl.constexpr,
    num_heads_k: tl.constexpr,
    headdim: tl.constexpr,
    page_block_size: tl.constexpr,
    blocks_per_batch,
):
    # 维度索引
    pid_b = tl.program_id(0)  # Batch index
    pid_h = tl.program_id(1)  # Head index

    # GQA Support: 映射 Query Head 到 KV Head
    kv_head = pid_h * num_heads_k // num_heads

    # 加载实际的 KV 序列长度
    seq_len = tl.load(cache_seqlens_ptr + pid_b).to(tl.int32)

    # 维度偏移量 [0, 1, ..., headdim-1]
    offs_d = tl.arange(0, headdim)

    # Online Softmax 累加器
    acc = tl.zeros([headdim], dtype=tl.float32)
    l_i = 0.0
    m_i = float('-inf')
    scale = 1.0 / tl.sqrt(float(headdim))

    # === 性能瓶颈：串行遍历整个序列 ===
    # 不使用 Block 并行，而是用单个 Block 串行循环处理所有 Token
    t = 0
    while t < seq_len:
        # 性能瓶颈：每次循环都重新加载 Q，增加显存压力
        q = tl.load(q_ptr + pid_b * num_heads * headdim + pid_h * headdim + offs_d).to(tl.float32)
        q = q * scale

        # Paged KV 映射逻辑
        page_idx = t // page_block_size
        page_off = t % page_block_size

        # 查表获取物理 Block 索引
        # blocks_per_batch 是计算出来的步长
        phys_block = tl.load(block_table_ptr + pid_b * blocks_per_batch + page_idx)

        # 计算 K 和 V 的物理地址
        # Layout: (num_blocks, page_block_size, num_heads_k, headdim)
        kv_base = phys_block * page_block_size * num_heads_k * headdim + \
                  page_off * num_heads_k * headdim + \
                  kv_head * headdim
        
        # 加载 K 和 V 向量
        k = tl.load(k_cache_ptr + kv_base + offs_d).to(tl.float32)
        v = tl.load(v_cache_ptr + kv_base + offs_d).to(tl.float32)

        # Attention 计算
        s = tl.sum(q * k)  # 点积
        
        # Online Softmax 更新
        m_new = tl.maximum(m_i, s)
        p = tl.exp(s - m_new)
        alpha = tl.exp(m_i - m_new)
        
        acc = acc * alpha + p * v
        l_i = l_i * alpha + p
        m_i = m_new

        t += 1

    # 写回结果
    # 这里没有处理 l_i 为 0 的边界情况，但测试数据 seq_len 通常很大
    out = acc / l_i
    tl.store(output_ptr + pid_b * num_heads * headdim + pid_h * headdim + offs_d, out)

def run_kernel(
    q, k_cache_paged, v_cache_paged, output,
    cache_seqlens, block_table,
    batch_size, seqlen_k, seqlen_q, num_heads, num_heads_k, headdim,
    page_block_size, num_blocks, causal,
):
    # 计算每个 batch 对应的 block_table 行宽
    blocks_per_batch = num_blocks // batch_size
    
    # 启动配置：每个 Head 一个 Block
    # 总 Block 数 = batch_size * num_heads (最大 128个)，并行度极低
    grid = (batch_size, num_heads)
    
    slow_decode_kernel[grid](
        q, k_cache_paged, v_cache_paged, output,
        cache_seqlens, block_table,
        num_heads=num_heads,
        num_heads_k=num_heads_k,
        headdim=headdim,
        page_block_size=page_block_size,
        blocks_per_batch=blocks_per_batch,
        num_warps=1,    # 性能瓶颈：仅使用 1 个 warp，限制计算吞吐
        num_stages=1,   # 性能瓶颈：禁用流水线并行
    )